昨天把 Phase 22 拆成三段:
建構 CRT → 求餘數 → INT8 GEMM → CRT 重建
今天只看第一段。
問題很小:
這 7 個質數跟那些
f_i,到底是怎麼來的?
我們先不碰矩陣乘法,也不碰 Tensor Core。今天只要做到一件事:給你一組餘數,kernel 能把原本的大整數拼回來。
這件事如果單獨不對,後面 GEMM 寫得再漂亮都沒用。
Phase 22 用這 7 個質數:
p = 127, 113, 109, 107, 103, 101, 97
第一個條件是:全部小於 128。
因為我們要把餘數塞進 int8_t,然後丟給 INT8 Tensor Core。int8_t 的有號範圍是:
-128 ~ 127
餘數本身會整理成 [0, p):
0 ~ 126
0 ~ 112
...
0 ~ 96
這些數字都能安全放進 int8_t。
這裡有一個小坑:int8_t 是有號型別,但我們存的餘數都小於 127,所以直接 cast 不會變號。也就是說:
d[...] = (int8_t)rem;
只要 rem 在 0..126,它在 Tensor Core 眼裡還是同一個正數。
不要選 131 這種質數。數學上可以,硬體資料型別不可以。
M 是我們能重建的整數空間7 個質數兩兩互質,所以總模數是它們的乘積:
M = 127 × 113 × 109 × 107 × 103 × 101 × 97
= 168897325606883
這是大約 1.69e14,約 47 bits。
CRT 保證:在 [0, M) 之間,每一組 7 個餘數只對應一個整數。
但 GEMM 的內積可能是負的,所以我們最後會把 [0, M) 折成以 0 為中心:
if x > M/2:
x = x - M
折完就是大約:
[-M/2, M/2]
所以 7 個質數真正能保護的是「有號整數結果不要超過這個範圍」。
對一個 tile 的內積來說,粗略上界是:
threshold ≈ K × max|A_slice| × max|B_slice|
只要 threshold < M/2,重建就不會繞錯圈。
這就是為什麼程式裡不只存完整的 M,也存前綴乘積:
d_M_arr[1] = 127
d_M_arr[2] = 14351
d_M_arr[3] = 1564259
d_M_arr[4] = 167375713
d_M_arr[5] = 17239698439
d_M_arr[6] = 1741209542339
d_M_arr[7] = 168897325606883
如果某個 tile 的 threshold 很小,理論上可以只用前 3 或 4 個質數。這就是 nl:
nl = number of limbs
nl 越小,INT8 GEMM 次數越少,CRT 重建也越便宜。
現在 Phase 22 的 fast path 固定用 7 個質數,因為它要走最直接、最穩的路徑。舊的 dynamic path 仍然保留 mA、mB、kDynamicNL 那套邏輯,讓每個 tile 可以自己選 nl。
f_i 怎麼算Day 5 用 3,5,7 拼過 23。今天換成 Phase 22 的質數。
對每個質數 p_i,先算:
M_i = M / p_i
然後找一個逆元 y_i:
M_i × y_i ≡ 1 (mod p_i)
最後:
f_i = M_i × y_i
f_i 的特性是:
f_i ≡ 1 (mod p_i)
f_i ≡ 0 (mod 其他 p_j)
所以它像一塊只對自己的質數有反應的積木。
重建時:
x = Σ residue_i × f_i mod M
每個 residue 只會在自己的模數上留下正確訊號,在其他模數上都是 0。
Phase 22 完整 7 質數的係數是:
| p | f_i |
|---|---|
| 127 | 147618922380819 |
| 113 | 112099994871825 |
| 109 | 134807957135769 |
| 107 | 34726552928518 |
| 103 | 96747011755399 |
| 101 | 130435558389474 |
| 97 | 19153304965729 |
這些數字在程式裡放在 constant memory:
__constant__ uint64_t d_coeffs_flat[29] = {
...
147618922380819,
112099994871825,
134807957135769,
34726552928518,
96747011755399,
130435558389474,
19153304965729
};
為什麼是 flat,不是 coeffs[8][7]?
因為 nl 可能是 1 到 7。每個 nl 都有自己的 M,所以係數也不一樣。d_coeff_offsets 用來告訴 kernel:
M = d_M_arr[nl];
off = d_coeff_offsets[nl];
cp = d_coeffs_flat[off + p];
如果 nl = 7,offset 是 22,後面那 7 個就是完整 Phase 22 的係數。
先想像我們已經有某個元素的 7 個餘數:
r[0] = x mod 127
r[1] = x mod 113
...
r[6] = x mod 97
最小的 reconstruct kernel 只要做:
uint64_t sum = 0;
for (int p = 0; p < nl; ++p) {
sum = (sum + r[p] * coeff[p]) % M;
}
int64_t signed_x = (sum > M / 2) ? sum - M : sum;
這就是 CRT 的全部。
在 Phase 22 fused kernel 裡,這段被塞進 crt_pass_kernel_body() 的尾端。它不是讀一個 r[p] array,而是直接拿 INT8 GEMM 的 accumulator:
v = acc % prime[p];
weighted_sum += v * coeff[p];
weighted_sum %= M;
然後:
if (weighted_sum > M / 2) signed_x = weighted_sum - M;
C[i,j] += signed_x * inv;
這裡的 inv 是 Ozaki 那個 pass 的縮放,例如 2^-34。
% 的第一個硬體坑Day 7 開始會碰到一批硬體跟指令層的坑。
第一個就是 %。
寫成:
v = acc % p;
很直覺,但 GPU 上整數除法不是便宜指令。p 又不是 2 的次方,不能靠 bit mask 解掉。
所以後來 Phase 22 fast path 對固定 7 質數做了 magic division:
q = __umulhi(x, magic[p]);
r = x - q * p;
if (r >= p) r -= p;
這段的前提是 x 的範圍已知。程式裡註解寫得很白:
// x in [0, 126*126*2048]
因為一個 INT8 GEMM accumulator 最多大約是:
126 × 126 × K
在 2048 的 K 下還放得進 uint32_t。
完整 M 的 modulo 也有類似優化,nl = 7 時可以用 Barrett reduction:
M = 168897325606883
mu = floor(2^64 / M) = 109218
Day 7 不需要把這些優化寫進最小 kernel。
但要先知道:% p 是正確性的好起點,不是最後的速度答案。
[0, p)第二個坑是 C/C++ 的 % 對負數不是數學課的 mod。
例如:
-3 % 127 == -3
但 CRT 要的是:
124
所以程式裡每次取完餘數都要做:
int32_t rem = iv % p;
if (rem < 0) rem += p;
這一行看起來無聊,但少了它,所有負數 slice 都會壞。
而且這個 bug 很陰險。隨機矩陣如果正負平均,你會看到某些元素差很多,某些又看起來還好。因為只有落到負餘數的那批路徑會錯。
Day 8 的 precompute_modulo_hl_A/B 也會做同一件事。
CRT 表很小:
7 個 primes
8 個 M 前綴
29 個 coefficients
8 個 offsets
這些資料所有 thread 都會反覆讀,而且整個執行期間不變。
所以程式放在:
__constant__ int d_primes[7];
__constant__ uint64_t d_M_arr[8];
__constant__ uint64_t d_coeffs_flat[29];
__constant__ int d_coeff_offsets[8];
對 GPU 來說,constant memory 適合這種小表。
原因有兩個。
第一,這些表在 kernel 執行期間不會變。每個 thread 都讀同一批 prime、M、offset、coefficient,所以沒有必要把它們塞進一般 global memory,然後每次都從 DRAM 慢慢撈。
第二,warp 裡很多 thread 會在同一個時間讀同一個表格位置。constant memory 對這種「所有 thread 廣播同一個常數」的模式很友善,cache 行為也比較漂亮。
這跟把 A、B 餘數矩陣放 global memory 是完全不同的事。A、B 是大資料,而且每個 tile、每個 thread 讀的位置都不一樣,所以會照 tile 搬進 shared memory。CRT 表是小常數,放 constant memory 就好。
不過這裡有幾個實作上很容易踩到的坑。
第一個坑是忘記初始化,或初始化順序錯了。
__constant__ 變數不是你在 host 端宣告完就自動有值。程式啟動時要明確用 cudaMemcpyToSymbol 把 host 端算好的表拷進去:
cudaMemcpyToSymbol(d_primes, h_primes, sizeof(h_primes));
cudaMemcpyToSymbol(d_M_arr, h_M_arr, sizeof(h_M_arr));
cudaMemcpyToSymbol(d_coeffs_flat, h_coeffs_flat, sizeof(h_coeffs_flat));
cudaMemcpyToSymbol(d_coeff_offsets, h_coeff_offsets, sizeof(h_coeff_offsets));
如果少拷一張表,kernel 不一定會直接 crash。更常見的狀況是結果看起來像 CRT 算錯,實際上只是 device 端讀到舊值或零值。
所以 debug 時我會先做一顆很小的檢查 kernel,只把 constant memory 裡的值複製回一個 debug buffer,確認 device 端看到的表和 host 端完全一樣。這一步很土,但很有效。
第二個坑是 d_coeff_offsets 算錯。
這張表的意義是「某個 nl 的係數從 flat array 的哪裡開始」。如果前面幾段係數數量算錯一格,nl = 7 讀到的就不會是那 7 個完整係數。
比較保險的檢查方式是列出每個 nl 的區間:
coeffs_flat[0] 是保留值
nl = 1: coeffs_flat[1..1]
nl = 2: coeffs_flat[2..3]
nl = 3: coeffs_flat[4..6]
...
nl = 7: coeffs_flat[22..28]
對應到程式裡就是:
__constant__ int d_coeff_offsets[8] = {0, 1, 2, 4, 7, 11, 16, 22};
這也解釋了為什麼完整 7 質數的 offset 是 22。扣掉 index 0 的保留值後,前面已經放了:
1 + 2 + 3 + 4 + 5 + 6 = 21
所以第 7 組從保留值後第 21 個係數的下一格開始,也就是 index 22。
第三個坑是型別寬度。
M 大約是 1.69e14,已經遠遠超過 uint32_t。M、f_i、weighted_sum 都要用 64-bit:
uint64_t M = d_M_arr[nl];
uint64_t coeff = d_coeffs_flat[off + p];
uint64_t weighted_sum = 0;
但 residue 本身可以很小:
uint32_t v = acc_mod_prime;
這樣分開寫,是為了讓資料意義清楚。小的餘數用小型別,大的重建空間用 64-bit。不要因為 prime 小於 128,就順手把整段 CRT 都寫成 32-bit。
第四個坑是 nl 沒有防呆。
nl 合法範圍是 1 到 7。如果 dynamic path 算出來的 nl 是 0,或超過 7,讀 d_M_arr[nl] 和 d_coeff_offsets[nl] 就會越界。這種錯在 GPU 上常常不會給你漂亮的錯誤訊息,只會在最後答案裡變成奇怪的大數。
所以我會在 host 端先擋:
nl = std::max(1, std::min(nl, 7));
或至少在 debug build 裡加 assert。fast path 固定 nl = 7,風險比較小;dynamic path 才是更需要檢查的地方。
真正開發時,我不會一開始就拿整個 GEMM 來驗證 constant memory。比較好的順序是:
p_i、M、f_i。x,只測 CRT reconstruct。x,確認不同 thread 讀 constant memory 也沒問題。這樣切開之後,錯誤會比較好定位。否則你看到最後的 C 矩陣不對,很難分辨是質數表錯、offset 錯、負數餘數錯、還是 GEMM accumulator 本身就錯。
hand_over_doc 和 migration_assets 裡有很多很有用的 debug 紀錄,但它們不是同一個時間點的規格書。有些內容是歷史階段留下來的,有些是把 CRT/RNS 這個想法搬到別的情境時做的整理。
如果你看到舊文件裡出現:
251, 241, 239, ...
那不是現在 Phase 22 signed INT8 路徑用的質數表。
現在這條路用的是:
127, 113, 109, 107, 103, 101, 97
原因不是數學上 251 不能用,而是這條 kernel 的資料路徑不允許。
現在的 Phase 22 signed INT8 路徑有三個硬限制:
int8_t。251, 241, 239, ... 這種質數表比較適合「餘數存在更大的型別裡」的版本。它們的乘積可以提供很大的重建範圍,但餘數本身可能超過 127,不能直接當 signed INT8 operand。
這也是為什麼讀舊文件時,不能只看 CRT 公式。CRT 公式本身可能都對,但硬體路徑不一樣,表就不能照搬。
我後來整理這段時,會把每份資料分成三種層級:
概念層:CRT 怎麼把餘數拼回整數
規格層:Phase 22 目前使用哪些 prime、M、f_i
實作層:CUDA kernel 實際怎麼取餘數、存 int8、重建
概念層可以跨版本參考。規格層和實作層就要很小心,因為一換 phase、一換資料型別,限制就不一樣。
debug 時最怕的不是完全不能跑,而是「看起來很像有跑,但數字慢慢歪掉」。我會優先查這幾件事:
127, 113, 109, 107, 103, 101, 97。M 是否一樣。f_i 是否是用同一組 M 和同一組 prime 算出來的。d_coeff_offsets[nl] 是否指到正確區間。[0, p)。M/2。其中第 3 點特別容易忽略。f_i 不是只跟單一質數有關,它跟整組 prime 的乘積 M 有關。你不能把舊表的 prime 和新表的 coefficient 混在一起用。
最小的驗證方式是不要碰矩陣,先固定幾個值:
x = 0
x = 1
x = -1
x = 123456789
x = -123456789
x = M/2 - 1
x = -(M/2 - 1)
每個 x 都走同一條流程:
CPU 算 residue
GPU reconstruct
copy 回 host
比對 signed_x
這些都過了,才有資格說 CRT 表和 reconstruct kernel 大致正確。
接著再用一個很小的矩陣測 GEMM,例如 2 × 2 乘 2 × 2,而且元素刻意放正數和負數:
A = [ 3, -5
-7, 11 ]
B = [ -13, 17
19, -23 ]
這種測試比大型 random matrix 更有用,因為每一個錯誤都能手算追蹤。大型 random matrix 適合最後壓測,不適合第一輪 debug。
我另外把建置和 debug 的坑整理在 Blog/debug_notes.md。真的要跑 repo 時,可以先看那篇,避免把舊 Phase 24、舊質數表、舊效能數字混進現在的 Phase 22。
這篇 Day 7 先把 CRT 表本身講清楚。後面 Day 8 進入 precompute_modulo_hl_A/B 時,會開始看到這些表怎麼被拿去產生 7 張 INT8 餘數矩陣。
Day 7 的產出不是快。
今天只要有這三個東西:
p_i、M、f_i。x,先算 7 個餘數,再讓 kernel 拼回 x。測試可以先不用矩陣。
例如:
x = -123456789
for each p:
r_p = x mod p, folded to [0, p)
kernel reconstruct(r_p) → -123456789
這個過了,再把輸入換成小矩陣的每個元素。
如果小整數都拼不回來,就不要往 GEMM 走。
今天可以帶走三句:
int8_t,所以不能選太大的模數。M 是可重建範圍。 K × max|A| × max|B| 不能超過 CRT 能保護的整數空間。f_i 是重建積木。 每個 f_i 只在自己的質數上等於 1,在其他質數上等於 0。明天進第二段:把 FP64 矩陣真的變成 7 張 INT8 餘數矩陣,然後先用最笨的 INT8 GEMM 對答案。